2.5條梧曾經說過
前幾天我們一直在想一件事:
能不能用 INT8 Tensor Core 做出比 FP64 更快的高精度 GEMM?
原生的硬體贏了
但 額 人家也知道這個方法 NV也不是吃白飯的
速度上 人家的int8 mma利用率高多了 呵呵呵
所以呢
如果目標只是「輸入 double、輸出 double,然後速度贏 cuBLAS」,這件事沒有想像中簡單。cuBLAS 本來就是 NVIDIA 自己長年優化的 GEMM library,尤其現在連 emulation 路線也做得很成熟。
所以
今天來寫 DoubleDouble
昨天我們其實很努力想把它做快。
但 profile 看完以後,現實也很明顯:要直接靠現在這條路去贏 cuBLAS,很難。不是說 INT8 GEMM 不能跑,而是整個流程裡面還有太多 GEMM 前後的東西。你把 MMA kernel 再磨一點,可能有幫助,但它不是現在最大的洞。
所以今天換方向。
速度先不贏,那可不可以把精度往上移?
今天就來寫 DoubleDouble。
DoubleDouble 的格式很簡單:
struct DoubleDouble {
double hi;
double lo;
};
數值語意是:
value = hi + lo
一個 double 大約有 53-bit significand。DoubleDouble 用兩個不重疊的 double 去放一個數,所以有效精度大約是:
53 + 53 = 106 bits
這裡不是要說它是某個硬體原生格式。它不是。
這次要做的事情也不是「宣告一個 struct 就結束」。真正要確認的是:
DoubleDouble 的 lo 有沒有真的進入 residue
residue GEMM 後有沒有被保留下來
CRT reconstruction 後有沒有回到 DoubleDouble output
只要中間任何一步把 hi + lo 直接變回普通 double,這件事就失敗了。
這次最關鍵的地方其實在 GEMM 前面。
不能這樣寫:
double x = dd.hi + dd.lo;
這行會把很多 lo 直接吃掉。
所以我們現在做的是把 DoubleDouble 轉成整數表示。
概念上是:
DoubleDouble(hi, lo)
-> 乘上一個 2^s scale
-> 拆成整數 words
目前單一元素用:
uint64_t words[3];
也就是 192-bit 的暫存整數表示。這裡只是單一 input element 的 scaled integer,不是最後 CRT reconstruction 的總空間。最後重建會用更寬的 UInt384。
轉換的核心是對 hi 和 lo 分別做:
frexp(x) -> mantissa / exponent
shift = exponent - 53 + scale
mantissa 放進 uint64_t words
正負號也要一起處理。hi 是負數時,整個 scaled integer 要用負的 residue 表示;lo 不能被當成獨立的普通 double 結果丟掉。
這段的重點不是浮點加法,而是:
先把 DoubleDouble 的有效位轉成一個整數,再進入 modulo 世界。
這樣後面每個 modulus 看到的才是同一個高精度整數。
一開始最直覺的寫法會變成:
for each modulus:
DoubleDouble -> integer
integer -> residue
這樣會讓同一個元素重複做很多次 frexp、round、shift。
現在改成:
for each element:
DoubleDouble -> words
for each modulus:
words -> residue
也就是每個元素只做一次 DoubleDouble 拆解。
words -> residue 的做法是把 64-bit word 再拆成 32-bit chunk,從高位往低位掃:
rem = (rem * 2^32 + chunk) mod p
最後把 residue 轉成 balanced int8:
if rem > p / 2:
rem -= p
所以 residue matrix 裡面放的不是原始數字,而是每個 modulus 下的代表值。
這裡也解釋了為什麼我們可以用接近 8-bit 的 modulus。像 p=256 時,255 可以用 -1 表示,因為在 mod 256 裡:
-1 == 255
Tensor Core 吃的是 signed int8,但數學上我們用 balanced residue 解釋它。
DoubleDouble 不是 53-bit target。
它大約是 106-bit target。
GEMM 裡每個輸出元素是 K 個乘積累加,所以 CRT 空間要估:
required_crt_bits = 2 * input_bits + ceil(log2(K))
對 DoubleDouble:
input_bits = 106
所以:
| case | required bits | selected modulus count | CRT bits |
|---|---|---|---|
| K=128 | 219 | 29 | 約 220.1 |
| K=4096 | 224 | 30 | 約 226.9 |
這也是為什麼 20 個 modulus 不夠。
20 個 modulus 大約只有:
155 bits
對 DoubleDouble 的乘加空間來說太小。
目前表裡可用的 modulus 到 49 個,整體 bit budget 大約 341.9 bits。DoubleDouble 不需要全拿;K=4096 目前約 30 個就夠。但重點是這個表要能往上選,而不是停在舊的 20 個。
如果使用者硬指定:
DoubleDouble
K = 4096
num_moduli = 20
現在正確行為是直接拒絕:
Insufficient CRT bits for DoubleDouble
這種拒絕很重要。高精度程式不能「算完再說」,因為 CRT 空間不夠時,它不是誤差大一點而已,而是會 wrap 到另一個整數。
產生 residue matrix 以後,資料長這樣:
A_mod[p][m, k]
B_mod[p][k, n]
對每個 modulus 做:
C_mod[p] = A_mod[p] * B_mod[p]
A_mod 和 B_mod 是 int8,accumulator 是 int32。
這裡的 kernel contract 很單純:
input: int8 residue matrix
output: int32 residue GEMM result
高精度不在單一 GEMM 裡完成。
高精度是靠:
多個 modulus 的 int32 GEMM result
再一起做 CRT reconstruction
所以這次我沒有把文章重點放在 MMA 指令。昨天已經看過 MMA 和 profile 了,今天真正新增的是 DoubleDouble 怎麼進入這條 residue pipeline。
C_mod[p] 算完以後,要把多個 modulus 的結果拼回同一個整數。
DoubleDouble K=4096 需要大約 224-bit CRT 空間,所以舊的 192-bit reconstruction 不夠。
這次用:
UInt384
重建時先算:
M_total = product(p_i)
M_i = M_total / p_i
inv_i = inverse(M_i mod p_i)
每個輸出元素再做:
result = sum C_mod[i] * inv_i * M_i
result = result mod M_total
如果:
result > M_total / 2
就把它視為負數區間。
這裡有一個實作細節:result 是很多個 M_i * term 加起來,理論上要做 mod M_total。目前因為項數有限,實作上用重複 subtract M_total 的方式把它壓回範圍內。這不是最快,但對這個 correctness path 夠直接,也比較好 debug。
重建完得到的是一個寬整數,還要乘回 scale。
如果這裡只轉成普通 double,前面保住的低位又會不見。
所以現在的做法是把 UInt384 拆成 32-bit chunk,逐段變成 double term,再用 DoubleDouble addition 累加:
for each 32-bit chunk:
term = chunk * 2^(chunk_offset) * scale
dd_acc += term
最後得到:
C.hi
C.lo
這段很重要,因為它讓 CRT reconstruction 的結果不是只落回 53-bit,而是回到 DoubleDouble 表示。
這種東西不能只拿 random matrix 測。
random input 很容易讓 FP64 看起來也還可以,或者讓錯誤被平均掉。
這次最有用的測試是故意做 low-bit cancellation。
設定:
A = +/-1 + 2^-60
B = 1
K = 128
FP64 collapse 後只看到 +1 和 -1,加起來是:
0
但 DoubleDouble 應該留下每一項的 2^-60:
128 * 2^-60 = 1.110223e-16
實測:
FP64 collapsed: 0
DoubleDouble: 1.110223e-16
oracle: 1.110223e-16
這個 case 很直接。它不是在看一般誤差漂不漂亮,而是在問:
FP64 看不到的那一段,我們是不是真的保存下來了?
答案是有。
目前 128x128x128 測了三組:
| case | DoubleDouble max_rel | FP64 collapsed max_rel |
|---|---|---|
| cancellation_lo_only | 0 | 1.0 |
| constant_lo_1e-20 | 約 1.64e-33 | 約 2.00e-20 |
| varied_hi_lo | 約 4.04e-27 | 約 7.92e-12 |
這張表的重點是對照組。
同一份輸入,如果先 collapse 成 FP64,低位會掉;走 DoubleDouble residue path,低位可以留下。
所以這次可以比較有把握地說:
DoubleDouble correctness path 是成立的
但不能說:
高速 DoubleDouble GEMM 已經完成
這兩句要分開。
這次有把 stage timing 印出來。
128³ 測試裡,INT8 GEMM 本身大概只有:
0.3 - 0.4 ms
但整個 DoubleDouble engine 大約:
23 - 31 ms
主要時間在:
DoubleDouble serialization: 約 13 - 15 ms
UInt384 CRT reconstruction: 約 8 - 15 ms
這個結果很有用。
它告訴我們:下一步如果要變快,不是優先去改 MMA 指令。
因為 GEMM 那段已經不是最大塊。
真正該打的是:
DoubleDouble -> residue
residue result -> UInt384 CRT
host / device round trip
materialized residue matrix traffic
也就是說,下一步的 fused 版本應該長這樣:
load DoubleDouble tile
-> 在 GPU 上轉 scaled integer / residue
-> 對固定 modulus group 做 INT8 GEMM
-> accumulator 做 mod reduction
-> 做 partial CRT
-> 減少寫回 global memory 的中間結果
如果只是把 per-modulus GEMM 換成手寫 MMA,但 serialization 和 CRT 還留在 host,速度不會有本質變化。
DoubleDouble 目前用的是 106-bit target。
如果未來想再往上,重點不是重寫概念,而是把這幾個東西加寬:
input precision bits
modulus count
CRT reconstruction width
output representation
oracle test
以 113-bit target、K=4096 來估:
required_crt_bits = 2 * 113 + 12 = 238 bits
目前 modulus table 前 32 個大約有:
240.4 bits
所以從 CRT 容量來看,往 113-bit target 走是可以的。
但今天先不把它叫完成。今天完成的是 DoubleDouble path,這是第一個站得住腳的高精度版本。
今天不是再講 CRT 原理,也不是再講 MMA 怎麼排。
今天真正做出來的是:
DoubleDouble input
-> scaled integer words
-> multi-modulus residue
-> INT8 GEMM per modulus
-> UInt384 CRT reconstruction
-> DoubleDouble output
這條路證明了:
只要 modulus count 和 reconstruction width 補夠,
FP64 看不到的低位可以被保留下來。
profile 也告訴我們下一步該做哪裡:
不是先改 MMA,
而是把 DoubleDouble serialization 和 CRT reconstruction 搬回 GPU,
再做真正的 fused residue / CRT pipeline。
速度還沒贏。
但精度往上推這件事,現在有具體版本了。
AI 還是在寫slob 但是今天我有多花時間迭代一下 同時也是我覺得至少我們可以做點什麼吧
雖然還是不能讓Nvidia emu滿意
但沒關西 生在有AI的時代真是太好了
可以保證沒工作 沒未來 只有一堆 債 要還
唉...(雖然其實跟AI也沒關係啦 是時機不好而已 沒事的)